DivGrad

计算逐元素除法 \(Y = x_1 / x_2\) 的反向梯度。支持同形逐元素,以及 x1 / x2 广播(一侧尺寸更大)的场景。

同形时:

\[\begin{split}\begin{aligned} dx_1 &= \frac{dy}{x_2} \\ dx_2 &= -\frac{dy \cdot x_1}{x_2^{2}} \end{aligned}\end{split}\]

广播时,先在广播后的公共形状上按上式计算,再对较小输入一侧沿广播轴做 ReduceSum,使 dx1 / dx2 与对应输入同形。 算子内部按 x1_sizex2_size 比较自动分发:x1 更大、x2 更大或两者相等。

输入:
  • dy - 上游梯度地址

  • x1_data - 前向被除数地址

  • x2_data - 前向除数地址

  • params - Parameter 结构体地址,布局见下

  • core_mask - 核掩码(仅共享存储版本使用)

Parameter 定义:

 1typedef struct {
 2    void *tile_data0;       // 临时工作空间 0,长度至少 dy_size 个元素
 3    void *tile_data1;       // 临时工作空间 1,长度至少 dy_size 个元素
 4    void *tile_data2;       // 临时工作空间 2,长度至少 dy_size 个元素
 5    int *large_shape;       // 较大输入形状,长度 ndims
 6    int *small_shape;       // 较小输入形状,长度 ndims
 7    int *out_shape;         // 输出 / dy 形状,长度 ndims
 8    long long ndims;        // 维数
 9    long long dy_size;      // dy 元素个数
10    long long x1_size;      // x1 元素个数
11    long long x2_size;      // x2 元素个数
12    int *large_strides;     // 较大输入 strides,长度 ndims
13    int *small_strides;     // 较小输入 strides,长度 ndims
14    int *out_strides;       // 输出 strides,长度 ndims
15    int *large_multiples;   // 较大输入广播倍数,长度 ndims
16    int *small_multiples;   // 较小输入广播倍数,长度 ndims
17    int *indices;           // 广播临时索引,长度 ndims,须清零
18    int *x1_shape;          // x1 形状,长度 ndims
19    int *x2_shape;          // x2 形状,长度 ndims
20} Parameter;
输出:
  • dx1 - 对 x1 的梯度,形状与 x1 相同

  • dx2 - 对 x2 的梯度,形状与 x2 相同

支持平台:

FT78NE MT7004

备注

  • FT78NE 支持 fp32

  • MT7004 支持 fp16、fp32

  • tile_data0 / tile_data1 / tile_data2 各需至少 dy_size 个元素;fp32 为 dy_size * sizeof(float) 字节,fp16 为 dy_size * sizeof(float16) 字节

  • *_strides / *_multiples 由算子内部使用,调用方分配长度 ndimsint 数组即可;indices 长度 ndims,须预先清零

  • large_shape / small_shape / out_shapex1_sizex2_size 关系设置:较大侧对应 large_shapeout_shape 取较大形状

共享存储版本:

void hp_div_grad_s(float16 *dy, float16 *dx1, float16 *dx2, float16 *x1_data, float16 *x2_data, Parameter *params, int core_mask)
void fp_div_grad_s(float *dy, float *dx1, float *dx2, float *x1_data, float *x2_data, Parameter *params, int core_mask)

C调用示例:

 1// MT7004 示例(共享存储多核,DDR 地址)
 2void TestDivGradSMCFp32(int core_mask) {
 3    int core_id = get_core_id();
 4    int logic_core_id = GetLogicCoreId(core_mask, core_id);
 5    int core_num = GetCoreNum(core_mask);
 6    float *dy = (float *)0x81000000;
 7    float *dx1 = (float *)0x82000000;
 8    float *dx2 = (float *)0x83000000;
 9    float *x1_data = (float *)0x84000000;
10    float *x2_data = (float *)0x85000000;
11    float *tile0 = (float *)0x86000000;
12    float *tile1 = (float *)0x87000000;
13    float *tile2 = (float *)0x88000000;
14    int *strides_base = (int *)0x8B000000;
15    int x1_shape[4] = {4, 4, 4, 4};
16    int x2_shape[4] = {4, 4, 4, 4};
17    Parameter params;
18    if (logic_core_id == 0) {
19        params.tile_data0 = tile0;
20        params.tile_data1 = tile1;
21        params.tile_data2 = tile2;
22        params.large_shape = x1_shape;
23        params.small_shape = x2_shape;
24        params.out_shape = x1_shape;
25        params.ndims = 4;
26        params.dy_size = 256;
27        params.x1_size = 256;
28        params.x2_size = 256;
29        params.large_strides = strides_base;
30        params.small_strides = strides_base + 4;
31        params.out_strides = strides_base + 8;
32        params.large_multiples = strides_base + 12;
33        params.small_multiples = strides_base + 16;
34        params.indices = strides_base + 20;
35        params.x1_shape = x1_shape;
36        params.x2_shape = x2_shape;
37        memset(params.indices, 0, 4 * sizeof(int));
38    }
39    sys_bar(0, core_num);
40    fp_div_grad_s(dy, dx1, dx2, x1_data, x2_data, &params, core_mask);
41}
42
43void main() {
44    int core_mask = 0b1111;
45    TestDivGradSMCFp32(core_mask);
46}

私有存储版本:

void hp_div_grad_p(float16 *dy, float16 *dx1, float16 *dx2, float16 *x1_data, float16 *x2_data, Parameter *params)
void fp_div_grad_p(float *dy, float *dx1, float *dx2, float *x1_data, float *x2_data, Parameter *params)

C调用示例:

 1// MT7004 示例(私有存储单核,AM 地址)
 2void TestDivGradAMFp32(void) {
 3    float *dy = (float *)0x10010000;
 4    float *dx1 = (float *)0x10016000;
 5    float *dx2 = (float *)0x10020000;
 6    float *x1_data = (float *)0x10026000;
 7    float *x2_data = (float *)0x10030000;
 8    float *tile0 = (float *)0x10036000;
 9    float *tile1 = (float *)0x10040000;
10    float *tile2 = (float *)0x10046000;
11    int *strides_base = (int *)0x10050000;
12    int x1_shape[4] = {4, 4, 4, 4};
13    int x2_shape[4] = {4, 4, 4, 4};
14    Parameter params;
15    params.tile_data0 = tile0;
16    params.tile_data1 = tile1;
17    params.tile_data2 = tile2;
18    params.large_shape = x1_shape;
19    params.small_shape = x2_shape;
20    params.out_shape = x1_shape;
21    params.ndims = 4;
22    params.dy_size = 256;
23    params.x1_size = 256;
24    params.x2_size = 256;
25    params.large_strides = strides_base;
26    params.small_strides = strides_base + 4;
27    params.out_strides = strides_base + 8;
28    params.large_multiples = strides_base + 12;
29    params.small_multiples = strides_base + 16;
30    params.indices = strides_base + 20;
31    params.x1_shape = x1_shape;
32    params.x2_shape = x2_shape;
33    memset(params.indices, 0, 4 * sizeof(int));
34    fp_div_grad_p(dy, dx1, dx2, x1_data, x2_data, &params);
35}
36
37void main() {
38    TestDivGradAMFp32();
39}